一种在 EAGLE-2 中提出的自适应草稿生成机制,通过 上下文感知的树结构,在保证生成质量的同时最大化投机解码的加速比。
静态草稿树(EAGLE‑1)固定每层分支数,但不同上下文的接受率差异巨大——简单位置可以多猜几个,困难位置猜多了全是废的。静态结构无法适应,浪费算力。
speculative_num_draft_tokens(默认 8),保证验证开销可控。以下用 极简版 EAGLE‑2 网络 和 数值示例 展示动态草稿树的完整推理过程。
⚡ 实际 EAGLE‑2 草稿模型通常为 2‑4 层 Transformer,远小于目标模型(如 7B/13B),因此生成开销极低。
📌 场景设置
max_nodes=8,top_k=3,max_depth=4草稿模型基于 “The cat” 的 hidden_states,并行预测前 3 个候选 token:
| 候选 token | 置信度 | 路径概率 |
|---|---|---|
sat | 0.85 | 0.85 |
slept | 0.60 | 0.60 |
jumped | 0.40 | 0.40 |
✅ 选择 Top‑3(top_k=3)全部保留,进入下一层。
对每个候选 token,继续预测下一个 token:
| 父节点 | 候选子节点 | 置信度 | 路径概率 (累积) |
|---|---|---|---|
sat | on | 0.90 | 0.85 × 0.90 = 0.765 |
sat | at | 0.50 | 0.85 × 0.50 = 0.425 |
slept | peacefully | 0.70 | 0.60 × 0.70 = 0.420 |
slept | on | 0.45 | 0.60 × 0.45 = 0.270 |
jumped | over | 0.55 | 0.40 × 0.55 = 0.220 |
jumped | up | 0.35 | 0.40 × 0.35 = 0.140 |
✅ 保留路径概率最高的 max_nodes=8 个节点(此处共有 6 个,全部保留,继续扩展)。
继续扩展,直到节点数达到 max_nodes=8,然后剪枝:
| 路径(累积概率排序) | 累积概率 | 状态 |
|---|---|---|
sat → on → the | 0.765 × 0.80 = 0.612 | ✅ 保留 (Top-1) |
sat → on → a | 0.765 × 0.30 = 0.230 | ✅ 保留 |
slept → peacefully → on | 0.420 × 0.75 = 0.315 | ✅ 保留 |
sat → at → the | 0.425 × 0.65 = 0.276 | ✅ 保留 |
slept → on → the | 0.270 × 0.70 = 0.189 | ✅ 保留 |
jumped → over → the | 0.220 × 0.60 = 0.132 | ✅ 保留 |
jumped → up → the | 0.140 × 0.40 = 0.056 | ❌ 剪枝 (概率最低) |
sat → on → mat | 0.765 × 0.10 = 0.077 | ❌ 剪枝 (概率低) |
✅ 最终保留 6 条路径(≤8),提交给目标模型验证。
将 6 条路径拼接成 batch,并行送入目标模型(如 LLaMA‑2 7B):
| 路径 | 验证结果 | 推进的 token |
|---|---|---|
sat → on → the | ✅ 全部接受 | sat, on, the |
slept → peacefully → on | ❌ 第1个 slept 即被拒 | 无 |
sat → at → the | ✅ sat 接受,❌ at 被拒 | sat |
slept → on → the | ❌ 第1个 slept 被拒 | 无 |
jumped → over → the | ❌ 第1个 jumped 被拒 | 无 |
sat → on → a | ✅ sat 接受,❌ on 被拒 | sat |
✅ 本轮实际推进:3 个 token( sat + on + the),仅用了 1 次大模型前向。
不是固定数量,而是树的总节点数固定,形状动态。 总节点数由 speculative_num_draft_tokens 控制,默认值为 8(部分模型默认 4)。
speculative_num_draft_tokens = speculative_num_steps × topk + 1 (近似,实际受剪枝影响)topk=4,num_steps=3 → 未剪枝前 1 + 4 + 16 + 64 = 85 个候选,但强制剪枝到 ≤8 个节点 再提交验证。
✅ 结论:EAGLE‑2 每次提交给目标模型验证的 草稿节点总数固定(如 8 个),但树的深度和每层宽度根据上下文实时变化。
| 角色 | 输入 | 输出 |
|---|---|---|
| 整体 | 当前已生成的 token 序列(prompt + 已生成部分) | 被目标模型接受的 token(可能 1~8 个),追加到序列中 |
| 草稿模型 | hidden_states(目标模型上一轮输出的隐藏状态) | 每个候选 token 的 置信度分数 + logits |
| 目标模型 | 草稿树中的所有 token 序列(拼接成 batch) | 每个位置的 接受/拒绝 决策,并顺带输出新 logits |
| 对比维度 | 静态草稿树 (EAGLE‑1) | 动态草稿树 (EAGLE‑2) |
|---|---|---|
| 树结构 | 固定深度、固定每层分支数 | 动态,根据置信度每层自适应 |
| 分支选择 | 全量扩展,无筛选 | 每层只保留 Top‑K 高置信度候选 |
| 剪枝策略 | 无(或简单按位置截断) | 综合概率剪枝,保留路径概率最高的 N 条 |
| 适应能力 | 弱,困难样本浪费算力 | 强,简单样本深搜,困难样本宽搜 |
| 加速比 | ~2‑3x | ~3‑5x(比 EAGLE‑1 提升 20‑40%) |
以下为 EAGLE‑2 动态草稿树构建的核心 Python 逻辑(基于 SafeAILab/EAGLE 简化示意)。
# ============================================================ # EAGLE‑2 动态草稿树构建(核心逻辑) # ============================================================ def eagle2_generate_tree( draft_model, hidden_states, # 目标模型上一轮输出的隐藏状态 token_ids, # 当前已生成的 token 序列 max_nodes = 8, # 总节点数上限(提交给目标模型验证的草稿数量) top_k = 4, # 每层保留的候选分支数 max_depth = 4 # 最大树深度 ): # 1. 初始化:根节点(当前已生成的位置) tree = [{ 'tokens': token_ids, 'hidden': hidden_states, 'prob': 1.0, # 路径累积概率 'depth': 0 }] # 2. 循环扩展,直到达到最大节点数或最大深度 while len(tree) < max_nodes and tree[-1]['depth'] < max_depth: current_nodes = [node for node in tree if node['depth'] == tree[-1]['depth']] if not current_nodes: break # 3. 并行前向:一次 forward 得到所有候选 token 的 logits 和置信度 logits, confidence = draft_model.forward(current_nodes) # 4. 对每个节点,选择 Top‑K 个最高置信度的 token 作为下一层分支 next_nodes = [] for idx, node in enumerate(current_nodes): top_indices = confidence[idx].topk(top_k).indices.tolist() for token_idx in top_indices: # 计算路径累积概率(用于后续剪枝) path_prob = node['prob'] * confidence[idx][token_idx].item() next_nodes.append({ 'tokens': node['tokens'] + [token_idx], 'hidden': logits[idx][token_idx], # 简化:实际为 hidden 更新 'prob': path_prob, 'depth': node['depth'] + 1 }) # 5. 剪枝:如果总节点数超过 max_nodes,只保留累积概率最高的 max_nodes 条路径 tree += next_nodes if len(tree) > max_nodes: tree.sort(key=lambda x: x['prob'], reverse=True) tree = tree[:max_nodes] # 6. 返回树结构(每个节点是一条候选路径) return tree # ============================================================ # 目标模型验证阶段(简略) # ============================================================ def verify_with_target(target_model, tree): # 将所有候选路径拼接成 batch,并行送入目标模型 batch = [node['tokens'] for node in tree] outputs = target_model.forward(batch) # 逐位置判断接受/拒绝(树验证 + 接受率计算) accepted_tokens = [] for i, tokens in enumerate(batch): if outputs[i].accepted: # 该路径被接受 accepted_tokens.append(tokens[-1]) # 取最后一个 token return accepted_tokens
| 参数 | 默认值 | 作用 |
|---|---|---|
speculative_num_draft_tokens | 8 | 目标模型一次验证的 总节点数(草稿数量上限),影响延迟和接受率 |
top_k | 4 | 每层扩展时保留的 候选分支数,越大树越宽 |
max_depth | 4 | 树的最大深度,限制最长依赖路径 |
confidence_threshold | 0.3 | 置信度截断阈值,低于此值的 token 直接丢弃(可选) |
💡 调参直觉:top_k 大 → 树更宽,覆盖更多候选,但总节点数固定(max_nodes=8),所以会压缩深度;max_depth 大 → 树更深,能生成更长的草稿,但浅层剪枝更激进。
top_k,所有分支均匀扩展max_nodes=8),算力利用最大化